Conversation
There was a problem hiding this comment.
Code Review
This pull request introduces dynamic fixed-m constants and safe bounds calculation based on KV sequence length, adds support for virtual K-centering, and enhances input validation and dtype safety checks in the Pallas flash attention kernel. It also expands the test suite to cover various edge cases, including per-Q-block fallback, batched isolation, and virtual K-centering. The review feedback focuses on optimizing the TPU kernel performance by replacing expensive dynamic integer division with optimized BlockSpec mapping and introducing a static use_k_centering flag to conditionally compile the centering logic at trace time.
3b9d7aa to
cb1e4d1
Compare
There was a problem hiding this comment.
A few things before this can go in though:
- If this lands on its own, fixed-m gets turned off for R>1 (
use_fixed_m = Falsein_ulysses_ring_custom_attention, it only comes back in #478). That breaks the recipe we ship today. Can we split the stack so each PR is safe to merge by itself? - Cutting the Ulysses gate from 213 to 113 feels more conservative than we need. With centered keys, the mass you can lose is bounded by 2^(ceil(U)-C-126), so a gate around C+116 already keeps the loss under 0.1%. Could you share fallback rates so we can pick the number?
mk[0]means something different now, but a 2Dmkstill gets silently broadcast. Can we just raise on that?- The V check never fired in my runs on WAN 2.2 (every call passed), and it only exists because C moved from 88 to 102. Is it worth the extra complexity?
- Please put back the comments that explain the Mosaic cliff and k-smoothing. They save the next person a lot of pain.
syhuang22
left a comment
There was a problem hiding this comment.
Some line-level notes to go with my comment above.
44dfe63 to
fc91941
Compare
d57775a to
678e806
Compare
678e806 to
07d7e72
Compare
b426abd to
73ab07f
Compare
7bef907 to
0254d18
Compare
05817c2 to
53905e1
Compare
…and safety fallbacks Implements exact fixed-m splash attention in Pallas on TPU. Instead of an online running max, a head / Q-block whose Cauchy-Schwarz bound max||q|| * max||k - k_bar|| is within a length-dependent window uses a precomputed m = ceil(bound) - C(N) (+ q . k_bar with virtual K-centering) and skips the per-block rescales; everything else keeps online softmax. - `get_fixed_m_constants(N)`: length-dependent recenter C(N) and safe bound that keep the f32 numerator/denominator from overflowing or flushing. - `fixed_m_dtype_is_safe`: rejects fp16/fp8 (exponent range too small for the recentered exponent); bf16 and fp32 pass. - V-magnitude gate: |V| <= DEFAULT_MAX_V_BOUND (256), otherwise online. - Per-(head, Q-block) dispatch in the kernel; the Ulysses and R=1 ring callers compile a uniform (all-fixed) and a hybrid kernel and pick one with `lax.cond` on `all_fixed`. Both kernels are compiled; the extra compile time / program size has not been measured. - Virtual K-centering: the kernel receives k_mean and adds q . k_bar to the fixed bound in VMEM (an MXU dot_general on an (Hkv, 8, d) tiled k_mean, precision=HIGHEST because the gate has no slack at its edge). K itself is never re-written in HBM. k_mean is padded to the lane-padded head_dim after `_pad_data_for_flash`, so head_dim < 128 works (regression test at d=64). Explicit metadata contracts on the fixed-m ring path ---------------------------------------------------- 1. Norm representation is declared, not inferred. The gate previously guessed whether `fixed_m_norms` were squared with `(qn.max() * mk.max()) < 1000.0`. Magnitude cannot answer that: legacy unsquared norms of (1000, 2) have a true bound of 2000, but read as squared give sqrt(2000) ~= 44.7, a ~45x under-estimate that admits fixed-m where it must fall back. Replaced by a REQUIRED `fixed_m_norms_squared` argument (no default; omitting it with fixed-m on raises). Tests use squared norms. 2. The V-safety predicate is required, not assumed. The V-magnitude and dtype verdict cannot be re-derived from one hop's Q/K, so the ring kernel raises unless `v_ok` is passed (`value=None` in the metadata helper gives v_ok = 0, fail-closed). `_ulysses_ring_custom_attention` (R > 1) computes it on the PRE-all-to-all shards -- dtype safety plus |V| <= 256 -- and reduces it with pmin over BOTH internal axes (ring, ulysses): before the a2a each device holds a sequence shard of every head, and the ring's accumulate-vs-LSE branch must be taken uniformly by every ppermute participant. Only Q and K go through the optimization_barrier that keeps XLA from duplicating their producer chains into the norm reductions; V is reduced outside it so the Q/K a2a does not wait on V's producer. 3. Norm shape is validated against per_q_block. Both gates compute `qn * mk[:, None]`, so a (num_heads,) array supplied with per_q_block=True broadcasts to (num_heads, num_heads) and pairs head j's query norm with head h's key norm; a sink head then inherits a small bound and the kernel evaluates exp2(large_logit - small_m) -> inf. The kernel now rejects the mismatch. Known limitations - On the uncentered R > 1 ring path the fixed bound can still flush the numerator for heads far below the bound; this is not measured here. - The ring kernel runs its own pmin over global-bound eligibility (added by this commit) unless the caller promises `uniform_fixed_m=True`; on the Ulysses+Ring path that pmin is redundant with the caller's (not measured). Tests - custom_splash_fixed_m_test.py (31): kernel unit tests with relative-L2 checks, k-norm (mk) shape and V-metadata fail-closed checks, integration through `_ulysses_attention` and the R=1 ring (uniform branch, hybrid branch via a sink head, d=64 -- the test re-derives the gate from the inputs and asserts it admits fixed-m, it does not observe the cond -- ragged KV tails with N % 8 == 0 and N % block_kv != 0 whose tail is an odd multiple of 8), and a U=2/R=2 check that a single V outlier on one pre-a2a device makes v_ok False on every device (ring kernel stubbed; needs >= 4 devices, explicit skip otherwise). - ring_fixed_m_test.py (18): 8 backend-independent contract tests (including the required fixed_m_norms_squared and v_ok arguments), 6 ring numerical tests vs an f32 reference (2 devices), 1 test that declared-unsquared and squared norms give the same result, and 3 tests showing the raw-K bound is unsound where the centered bound is not. - CI: to keep GitHub Actions short, 18 of these 49 tests (the kernel numerics grids and the multi-device integration tests) skip when GITHUB_ACTIONS=true. The 9 tests from main and the contract, dtype and V-metadata tests still run there. The new end_to_end/tpu/run_wan_stack_tests.sh runs all of them on a TPU VM; later PRs in this stack add their test files to it. Verified (final tree): TPU v6e-8 (jax 0.11.2, end_to_end/tpu/run_wan_stack_tests.sh): 49 passed, 6 subtests passed in 170.4s.
53905e1 to
4ec211c
Compare
Summary
Fixed-m splash attention kernel. When a provable bound on every logit (
‖q‖·‖k‖, optionally centred on the mean key k̄) fits the fp32 exponent range, the kernel uses that fixed maxminstead of an online running max. This removes the per-block max/rescale work from the softmax inner loop. When the bound can't be proven, it falls back to the standard online-softmax kernel (_lse_scan).What's in it
custom_splash_attention.py: fixed-m kernel with dynamic Q/KV bounds. The in-kernelq·k̄dot runs atPrecision.HIGHEST, and_splash_attention_forward_ringrejects 2D(2, num_q_heads)mkmetadata in favor of 3D(2, num_q_heads, num_q_blocks).fixed_m_norms_squaredis a required argument.value=Nonegivesv_ok=False(fail-closed)._custom_ring_attention_forwardvalidatesfixed_m_norms[0]shape againstper_q_block.< flash_min_seq_length) falls back todot_product, cross-attention remaps totokamax_flash, andheads % num_shards != 0raisesValueError.Known limitations
pminover global-bound eligibility unless the caller passesuniform_fixed_m=True.Tests
custom_splash_fixed_m_test.py(31) andring_fixed_m_test.py(18):mkshape and V-metadata fail-closed checks;fixed_m_norms_squaredandv_okcontract tests (ring_fixed_m_test.py);N%8==0,N%bkv≠0, odd-multiple-of-8 tail), and a U=2/R=2 pre-all-to-all V-outlier check;GITHUB_ACTIONS=trueto keep CI fast;end_to_end/tpu/run_wan_stack_tests.shruns all 49 on a TPU VM.run_wan_stack_tests.sh, 170.4 s).Stack: #477 → #478 → #479 → #488 → #491